import os
import pandas as pd
import matplotlib.pyplot as plt

def compare_results():
    # Specify the paths of your CSV files
    csv_files = [
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8n/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8s/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8m/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8l/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8x/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11n/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11s/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11m/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11l/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11x/results.csv"
    ]
    
    for file in csv_files:
        if os.path.exists(file):
            model_name = file.split('/')[-2]
            df = pd.read_csv(file)
            
            print(f"\nResults for {model_name}:")
            print(df[['epoch', 'metrics/accuracy_top1', 'metrics/accuracy_top5']].tail(1))
        else:
            print(f"File not found: {file}")

def plot_results():
    csv_files = [
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8n/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8s/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8m/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8l/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolov8x/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11n/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11s/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11m/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11l/results.csv",
        "C:/Users/caruso/Desktop/YOLO_Project/agrilus/yolo11x/results.csv"
    ]
    
    plt.figure(figsize=(10, 6))
    for file in csv_files:
        if os.path.exists(file):
            model_name = file.split('/')[-2]
            df = pd.read_csv(file)
            plt.plot(df['epoch'], df['metrics/accuracy_top1'], label=f'{model_name} Top-1 Accuracy')
            plt.plot(df['epoch'], df['metrics/accuracy_top5'], label=f'{model_name} Top-5 Accuracy', linestyle='--')
        else:
            print(f"File not found: {file}")
    
    plt.xlabel('Epoch')
    plt.ylabel('Accuracy')
    plt.title('Accuracy Comparison between Models')
    plt.legend()
    plt.grid(True)
    plt.show()

if __name__ == '__main__':
    compare_results()
    plot_results()
